Skip to content

概率、Softmax 与信息论 ​

标签
AI/ml/数学
字数
3699 字
阅读时间
15 分钟

上一层(03-线性代数与 GEMM)管的是张量怎么变换。这一层管的是模型输出的含义 —— logits 经 Softmax 变成分布,交叉熵对应负对数似然,KL 衡量分布差异。

随机变量与分布 ​

离散随机变量用概率质量函数,连续随机变量用概率密度函数:

pX(x)=P(X=x), ∑xpX(x)=1∫−∞∞pX(x)dx=1

连续变量在单点的概率为 0,区间概率由密度积分给出。

分布取值典型场景
Bernoulli0/1二分类、Dropout mask
Categorical1,…,V下一个 token 的分布
Uniform区间 / 有限集合初始化、随机采样
Gaussian实数初始化、噪声、近似分析

联合、边缘、条件 ​

P(X=x)=∑yP(X=x,Y=y)P(X=x∣Y=y)=P(X=x,Y=y)P(Y=y)

自回归语言模型用链式法则分解序列概率:

P(x1,x2,…,xT)=∏t=1TP(xt∣x1,…,xt−1)

这条式子解释了为什么训练时的 next-token prediction 与推理时逐 token 生成是同一个概率模型的两个过程。 训练在每个位置同时算这一项,推理按顺序一次算一项 —— 数学上是同一个分解式。

链式法则把序列概率分解成一连串条件概率 —— 这一条解释了训练与推理是同一个概率模型的两个过程。

   P(x₁, x₂, …, x_T) = ∏_{t=1}^{T} P(x_t | x₁, …, x_{t-1})

   ┌─ 训练(teacher forcing)───────────────────────────────┐
   │ 一次前向,在每个位置同时算这一项                          │
   │  位置 1   P(x₁)                                        │
   │  位置 2   P(x₂ | x₁)                                   │
   │  位置 3   P(x₃ | x₁ x₂)                                │
   │  …                                                     │
   │  └─ T 个位置并行算完 ──▶ 训练侧能吃满算力                 │
   └───────────────────────────────────────────────────────┘
   ┌─ 推理(自回归生成)────────────────────────────────────┐
   │ 按顺序一次算一项                                       │
   │  P(x₁) ─▶ 采样 x₁ ─▶ P(x₂ | x₁) ─▶ 采样 x₂ ─▶ …        │
   │  └─ 必须串行 ──▶ 推理成本由生成长度决定                    │
   └───────────────────────────────────────────────────────┘

   数学上是同一个分解式,差别只在「哪些项已经有值」。

   配套的一条:分布式训练隐含 i.i.d. 假设
     数据并行切分、梯度取平均、loss 汇总都建立在
     「各卡样本独立同分布」上。
     真实数据里的重复样本、分片偏差、序列相关性会破坏这个假设 ——
     这也是「按样本加权」而不是「简单平均」的原因。

独立与条件独立 ​

若 P(X,Y)=P(X)P(Y) 则独立。独立比「不相关」更强:协方差为 0 不一定独立。给定 Z 后条件独立写作 P(X,Y∣Z)=P(X∣Z)P(Y∣Z)。

分布式训练隐含了 i.i.d. 假设。 数据并行切分、梯度取平均、loss 汇总都建立在「各卡样本独立同分布」上。真实数据里的重复样本、分片偏差、序列相关性会破坏这个假设 —— 这也是「按样本加权」而不是「简单平均」的原因(见 02-分布式训练总论与显存账本 的加权均值算例)。

Bayes ​

P(X∣Y)=P(Y∣X)P(X)P(Y)

先验 × 似然 → 后验。常规 LLM 训练不直接算它,但它是概率推断、参数估计与不确定性建模的基础。

期望、方差、协方差 ​

E[X]=∑xxP(X=x)E[aX+bY]=aE[X]+bE[Y]

期望的线性性不要求变量独立 —— 这条在推导里常用。

Var(X)=E[(X−μ)2]=E[X2]−μ2

E[X2]−μ2 这个形式数值上不稳定(两个大数相减),正确做法见 01-数值计算与精度 的灾难性消减一节。

协方差与相关系数:

Cov(X,Y)=E[(X−μX)(Y−μY)],ρX,Y=Cov(X,Y)σXσY

协方差矩阵 Σ=E[(x−μ)(x−μ)⊤] 是对称半正定矩阵。PCA 对它做特征分解,找出方差最大的正交方向。

除以 n 还是 n−1 ​

s2=1n−1∑i=1n(xi−x¯)2

深度学习中「方差」是否除以 n 还是 n−1 取决于具体定义:

  • LayerNorm / BatchNorm 用总体方差形式(除以元素数 n)
  • 统计库的默认值可能是无偏形式(除以 n−1)

复现算子时要查 API 语义,不要只看名字。这两者在小 n 下差异可观,是「照着描述实现却对不上数值」的一个常见来源。

最大似然与负对数似然 ​

θ∗=arg⁡maxθ∏ipθ(yi∣xi)=arg⁡maxθ∑ilog⁡pθ(yi∣xi)

取对数把乘积变成求和 —— 既避免数值下溢,也便于求导。等价地最小化负对数似然:

LNLL=−∑ilog⁡pθ(yi∣xi)

这就是语言模型训练损失的定义本身。

Softmax ​

logits 不是概率 ​

模型最后一层输出 V 个任意实数 z=(z1,…,zV) —— 可为负,不需要和为 1。Softmax 把它们转成分布:

pi=ezi∑jezj

Softmax 对统一平移不敏感:

softmax(z+c)=softmax(z)

因为分子分母同乘 ec。这个性质正是数值稳定实现的依据 —— 既然平移不改变结果,那就平移到一个不会溢出的位置。

数值稳定的实现 ​

直接算 e1000 会溢出。令 m=maxjzj:

pi=ezi−m∑jezj−m

最大的指数输入变成 0,因此最大指数值为 1,其他都不大于 1。

python
def stable_softmax(x):
    m = max(x)
    exps = [exp(v - m) for v in x]
    return [v / sum(exps) for v in exps]

减最大值防的是正向溢出。很小的项仍可能下溢到 0 —— 这通常说明它相对最大项确实可忽略,但如果后续还要取对数,就不能走这条路径:应该直接用稳定的 log_softmax,而不是先变成 0 再 log(0)。

Softmax 对统一平移不敏感 —— 这是数值稳定实现的依据。

   softmax(z + c) = softmax(z)          因为分子分母同乘 e^c
        │
        └─ 既然平移不改变结果,那就平移到一个不会溢出的位置

   直接算 e^1000 会溢出,令 m = max_j z_j:
        p_i = e^{z_i − m} / Σ_j e^{z_j − m}
        │
        └─ 最大的指数输入变成 0 ⟹ 最大指数值为 1,其他都不大于 1

   def stable_softmax(x):
       m = max(x)
       exps = [exp(v - m) for v in x]
       return [v / sum(exps) for v in exps]

   两条边界
     减最大值防的是「正向溢出」
     很小的项仍可能下溢到 0 —— 这通常说明它相对最大项确实可忽略,
     但如果后续还要取对数,就不能走这条路径:
       应该直接用稳定的 log_softmax,而不是先变成 0 再 log(0)

   LogSumExp 的形式
     LSE(z) = log Σ e^{z_j} = m + log Σ e^{z_j − m}
     ⟹ log softmax(z)_i = z_i − LSE(z)
     └─ 交叉熵实现通常融合 log_softmax + NLLLoss:
        既减少中间张量与访存,也避开了「先 Softmax 再取 log」这条不稳定路径。
        这是「数值稳定与性能优化方向一致」的典型例子。

LogSumExp ​

LSE(z)=log⁡∑jezj=m+log⁡∑jezj−m

于是:

log⁡softmax(z)i=zi−LSE(z)

交叉熵实现通常融合 log_softmax + NLLLoss —— 既减少中间张量与访存,也避开了「先 Softmax 再取 log」这条不稳定路径。这是「数值稳定与性能优化方向一致」的典型例子。

温度 ​

pi(T)=ezi/T∑jezj/T
T效果
<1放大 logit 差异,分布更尖锐
>1缩小差异,分布更平坦
→0+趋近只选最大 logit(贪心)
→∞趋近均匀分布

温度改变的是采样分布,不改变模型权重。 实现上应先缩放 logits 再应用稳定 Softmax;非常小的 T 要谨慎处理(z/T 可能溢出)。

温度直接改变采样分布的形状。

   p_i(T) = e^{z_i / T} / Σ_j e^{z_j / T}

   T → 0⁺   ──▶  趋近只选最大 logit(贪心)
   T < 1    ──▶  放大 logit 差异,分布更尖锐
   T = 1    ──▶  模型的自然分布
   T > 1    ──▶  缩小差异,分布更平坦
   T → ∞    ──▶  趋近均匀分布

   └─ 温度改变的是采样分布,不改变模型权重。
      实现上应先缩放 logits 再应用稳定 Softmax;
      非常小的 T 要谨慎处理(z / T 可能溢出)。

   同一个分布上的三种切池方式
     Greedy            选最大概率 token
     Top-k             只在概率最高的 k 个中采样
     Top-p (nucleus)   取累计概率至少达 p 的最小候选集,再归一化采样
     └─ Top-p 的集合大小随分布尖锐程度自动变化 ——
        这是它成为多数聊天 API 默认的原因
     └─ 实现要点:通常先过滤 logits(把被排除项设为负无穷),
        再执行稳定 Softmax 与采样

   分布式推理下有一个容易漏掉的正确性问题
     若 vocabulary 被张量并行切分,global top-k / top-p
     需要跨设备聚合局部候选,或使用等价的分布式算法。
     └─ 各卡各自取局部 top-k 会得到与单卡不同的候选集。

交叉熵、熵与 KL ​

交叉熵 ​

H(q,p)=−∑iqilog⁡pi

真实标签是 one-hot(正确类别为 y)时:

L=−log⁡py

模型给正确类别的概率越高,损失越小;py→1 时损失为 0。

熵 ​

H(p)=−∑ipilog⁡pi

全部概率集中在一类时熵最小;V 类均匀分布时熵最大,为 log⁡V。对数底决定单位(自然对数 → nat,底 2 → bit),机器学习损失通常用自然对数。

KL 散度 ​

DKL(q‖p)=∑iqilog⁡qipi≥0

两分布相同时为 0,但通常不对称:

DKL(q‖p)≠DKL(p‖q)

所以 KL 不是严格意义的距离。

交叉熵可以分解:

H(q,p)=H(q)+DKL(q‖p)

训练时真实分布 q 固定,H(q) 与模型参数无关 —— 所以最小化交叉熵等价于最小化 KL 散度。这是「用交叉熵做损失」的理论依据。

KL 出现在知识蒸馏、分布匹配、RLHF/PPO 的约束项、投机解码的分析里。实现时要明确三件事:API 接收的是概率、log 概率还是 logits,KL 的方向(不对称!),以及沿哪个维度归约。

三个量的关系决定了「为什么用交叉熵做损失」。

   H(q, p)     = −Σ q_i log p_i                交叉熵
   H(q)        = −Σ q_i log q_i                熵
   D_KL(q ‖ p) = Σ q_i log(q_i / p_i) ≥ 0      KL 散度

        H(q, p) = H(q) + D_KL(q ‖ p)
        └─ 交叉熵 = 熵 + KL

   各自的边界情形
     概率全部集中在一类   熵最小
     V 类均匀分布         熵最大,为 log V
     两分布相同           KL = 0(但通常不对称,所以 KL 不是严格意义的距离)

   由此得到的推论
     训练时真实分布 q 固定 ⟹ H(q) 与模型参数无关
     ⟹ 最小化交叉熵 等价于 最小化 KL 散度
     └─ 这是「用交叉熵做损失」的理论依据

   真实标签是 one-hot(正确类别为 y)时,交叉熵退化成
     L = −log p_y
     └─ 模型给正确类别的概率越高,损失越小;p_y → 1 时损失为 0

   KL 在别处也反复出现:知识蒸馏、分布匹配、RLHF / PPO 的约束项、
   投机解码的分析。实现时要明确三件事:
     API 接收的是概率、log 概率还是 logits
     KL 的方向(不对称!)
     沿哪个维度归约

Perplexity ​

PPL=eL¯

(L¯ 为平均 token 负对数似然。)直觉上可理解为「模型每步面对的有效候选数」。

PPL 不能脱离评测设置横比

tokenizer、数据预处理、上下文长度、是否忽略特殊 token,都会显著改变 PPL 数值。只有词表与测试集相同才有比较意义 —— 这一点在 01-语言模型演进:N-gram 到 RNN 有更完整的讨论。

Top-k 与 Top-p ​

策略做法
Greedy选最大概率 token
Top-k只在概率最高的 k 个中采样
Top-p(nucleus)取累计概率至少达 p 的最小候选集,再归一化采样

Top-p 的集合大小随分布尖锐程度自动变化 —— 这是它成为多数聊天 API 默认的原因。两者的失效场景与参数交互见 04-采样参数。

实现要点:通常先过滤 logits(把被排除项设为负无穷),再执行稳定 Softmax 与采样。

分布式推理下的 top-k / top-p

若 vocabulary 被张量并行切分,global top-k / top-p 需要跨设备聚合局部候选,或使用等价的分布式算法。这是一个容易漏掉的正确性问题:各卡各自取局部 top-k 会得到与单卡不同的候选集。

「方差」除以 n 还是 n−1 取决于具体定义。

   总体方差(除以元素数 n)
     └─ LayerNorm / BatchNorm 用的是这一种

   无偏样本方差(除以 n−1)
     s² = (1/(n−1)) Σ (x_i − x̄)²
     └─ 统计库的默认值可能是这一种

   ┌─ 复现算子时要查 API 语义,不要只看名字。
   │  这两者在小 n 下差异可观,
   └─ 它是「照着描述实现却对不上数值」的一个常见来源。

   另一处同类陷阱:E[X²] − μ² 数值上不稳定
     Var(X) = E[(X − μ)²] = E[X²] − μ²
     └─ 右边是两个大数相减 —— 走的是灾难性消减那条路径
        正确做法见数值计算那篇的消减一节

   两条常被用到的性质
     期望的线性性:E[aX + bY] = aE[X] + bE[Y] —— 不要求独立
     独立比「不相关」更强:协方差为 0 不一定独立

相关 ​

参考 ​

贡献者 ​

文件历史 ​